Attention Cache and Sequence Parallelism

导言

KV cache 最容易造成的误解,是把所有名为 K/V 或 state 的张量都看成同一种“缓存”。事实上,训练时为反向传播保留的 K/V activation、推理时跨 decode step 存活的 KV cache、GDN 跨 token 改写的固定状态,生命周期和分布式处理都不同。

本文从一次 token 生成开始,逐对象解释 MHA、DSA 与 GDN 保存什么,再讨论 CP、USP 切开长序列后,临时 activation 和持久 cache 分别需要怎样的通信。

结论

先记住六个判断:

  1. KV cache 主要服务自回归推理。Prefill 一次构建 prompt 的逐层历史,decode 每步追加新 K/V;旧 Q 不缓存,因为未来不会再用它去发起查询。
  2. 训练通常没有跨 forward 持久增长的 KV cache。Teacher forcing 已知整段目标,一次并行形成所有 Q/K/V;K/V 只是当前 batch 的 activation,为 backward 保存或重算。
  3. MHA 保存逐 token、逐 KV head 的 K/V,cache 随上下文长度线性增长。
  4. DSA 仍有两套随长度增长的历史:Lightning Indexer 的低维 K/scale,以及 MLA 的 cKV/kR;Indexer 扫描全历史,主 MLA 只读取 global top-k。
  5. GDN 没有逐 token KV cache。它保存固定大小矩阵状态和有界短卷积状态;代价是历史被压进有限容量,而不是无损保留每个 token。
  6. CP/USP 的训练通信不能直接等同于 decode cache 分片。训练时 Q/K/V 都有长序列维;decode 的 Query 长度通常只有 1,若把历史 cache 按 token 分片,还要额外做 distributed softmax 或 global top-k。
![训练激活、KV cache 与递归状态的生命周期](https://pic.shaojiemike.top/shaojiemike/2026/07/ecdd1c05dba5e3bb12df0bf04b8a399e.png){ width=94% }
小黑认知示意图:训练 activation 是一次作业的临时工位,MHA cache 是不断增长的历史抽屉,GDN 则反复改写一块固定尺寸状态板。

先分清三种“保存”

为什么只缓存 K/V

在第 l 层、第 t 个位置,MHA 先把当前层输入 h_t^l∈R^D 投影成:

1
2
3
q_t^l = reshape(h_t^l W_Q^l)    [B,Hq,1,d]
k_t^l = reshape(h_t^l W_K^l) [B,Hkv,1,d]
v_t^l = reshape(h_t^l W_V^l) [B,Hkv,1,d]

对于 causal decoder,历史 token j≤tk_j^l,v_j^l 一旦算出,以后不会因为生成了新 token 而改变。第 t+1 步的新 Query 要再次读取它们,所以 cache 执行:

1
2
3
4
5
6
K_cache^l ← concat(K_cache^l, k_t^l, axis=time)
V_cache^l ← concat(V_cache^l, v_t^l, axis=time)

a_t^l = (q_t^l (K_cache^l)^T) / sqrt(d) + causal_mask
p_t^l = softmax(a_t^l, axis=history)
o_t^l = p_t^l V_cache^l

q_j^l 的工作在第 j 步已经结束;未来发起查询的是新 q_t^l,所以通常没有 Q_cache。Transformers 的通用动态 cache 也按层保存 shape 为 [B,Hkv,T,d] 的 K 和 V,并沿倒数第二维 append。[^hf-cache]

Prefill 与 decode

  • Prefill:输入 prompt 的 S 个 token 已知,可以用矩阵乘法一次形成所有层的 Q/K/V,并把各层 K/V 写入初始 cache。它不是逐 token Python 循环,但 causal mask 仍保证位置 i 看不到未来位置。
  • Decode:每步通常只有一个新 token。模型只为它计算新 Q/K/V,把新 K/V 追加到各层 cache;新 Q 读取从 prompt 到当前步的全部历史。
  • 每层独立保存:第 l 层的 K/V 来自该层输入,不能拿第 l-1 层 cache 代替。

以 BF16、L=32 层、B=1Hkv=32d=128T=8192 为例,忽略对齐和管理开销:

1
2
3
M_KV = 2 × L × B × Hkv × T × d × 2 bytes
= 4,294,967,296 bytes
= 4 GiB

这里第一个 2 是 K 与 V 两份,最后一个 2 bytes 是 BF16 元素大小。换成 GQA/MQA 会减少 Hkv,但 attention score 仍要覆盖全部历史 token。

训练为什么不同

训练使用 teacher forcing:目标序列在一次 forward 前已经全部给定。所有位置的 Q/K/V 可以并行计算,causal mask 只限制可见性。K/V 可能因 autograd 被保存到 backward,也可能由 activation checkpointing 重算,但它们通常在该次迭代结束后释放,不会在下一批数据上继续 append

同一个 use_cache 不代表训练也该开

Hugging Face 官方文档明确把 cache 限定为 inference;训练时启用可能引发非预期错误。[^hf-cache] 训练图需要梯度、dropout 和完整序列语义,生成 cache 的原地更新与跨 forward 生命周期并不匹配这些要求。

MHA:逐 token 保存 K/V

MHA 的对象账本如下。为了突出 cache,省略 MLP 和残差分支。

对象 公式 / 代码名 Shape 生产者 → 消费者 生命周期
当前层输入 h_t / hidden_states [B,1,D] 上一子层 → Q/K/V 投影 临时 activation
当前 Query q_t / query_states [B,Hq,1,d] WQ、reshape、RoPE → QK MatMul 临时 activation
当前 Key/Value k_t,v_t [B,Hkv,1,d] WK/WV、reshape、RoPE(K) → cache update 临时 activation
历史 Key/Value K_cache,V_cache [B,Hkv,T,d] 旧 cache + 当前 K/V → score/value MatMul 跨 decode step cache
分数与概率 a_t,p_t [B,Hq,1,T] QK MatMul、scale、mask、Softmax → PV MatMul 临时 activation
输出 o_t,y_t [B,Hq,1,d] → [B,1,D] PV MatMul、concat、WO → 下一子层 临时 activation

下面是完整的单层 decode 伪代码。repeat_kv 表示 GQA/MQA 为计算建立的逻辑 head view;MHA 中 Hkv=Hq,它是恒等操作。实际 kernel 可以不物化复制。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
def mha_decode(hidden_t, layer_cache, wq, wk, wv, wo, rope, scale):
batch, one, model_dim = hidden_t.shape
q = (hidden_t @ wq).view(batch, one, Hq, head_dim).transpose(1, 2)
k = (hidden_t @ wk).view(batch, one, Hkv, head_dim).transpose(1, 2)
v = (hidden_t @ wv).view(batch, one, Hkv, head_dim).transpose(1, 2)
q = rope(q, position=layer_cache.length)
k = rope(k, position=layer_cache.length)
k_all = concat([layer_cache.k, k], dim=2)
v_all = concat([layer_cache.v, v], dim=2)
layer_cache.k = k_all
layer_cache.v = v_all
k_for_q = repeat_kv(k_all, groups=Hq // Hkv)
v_for_q = repeat_kv(v_all, groups=Hq // Hkv)
scores = matmul(q, transpose(k_for_q, -1, -2)) * scale
probs = softmax(scores, dim=-1)
heads = matmul(probs, v_for_q)
merged = heads.transpose(1, 2).reshape(batch, one, model_dim)
return merged @ wo, layer_cache
![MHA prefill 与 decode 的 KV cache 逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/6a2ea7e9538f7c19a91101ef3301ce86.png){ width=98% }
物理语义图:蓝色圆点是当前 activation,双圈是跨 decode step cache,方块是具体算子;训练分支中的 K/V 只在本轮 forward/backward 存活。

DSA:两套随序列增长的压缩缓存

DSA 不是把 MHA cache 直接截成 top-k。DeepSeek-V3.2-Exp 在 MLA 前放置 Lightning Indexer:先用低维 Key 为每个 Query 选择历史位置,再让主 MLA 只读取这些位置。[^dsa-release]

对象 源码 / 公式名 Shape 生产者 → 消费者 生命周期
Indexer Query q_i [B,Hidx,Didx] 当前 hidden → Indexer Q 投影、RoPE、量化 临时 activation
Indexer Key k_j [B,T,Didx] FP8 + scale 每个历史 hidden → k_cache/k_scale_cache 随 T 增长的 cache
Indexer head 权重 w_i / weights [B,Hidx] 当前 hidden → 加权跨 head score 临时 activation
Index score I_j [B,T] FP8 dot、ReLU、跨 head 求和 → top-k 临时 activation
选择位置 J=TopK(I,k) [B,k] Indexer → 主 MLA gather / mask 临时控制对象
MLA latent cache cKV_j,kR_j [B,T,Rkv][B,T,DR] KV down projection、RMSNorm、RoPE → 主 MLA 随 T 增长的 cache
主 MLA 输出 y_t [B,D] 只聚合 J 中 latent → WUV/WO 临时 activation

固定 revision 的 Indexer 把低维 K 和量化 scale 物理写入两块 cache,再对全历史求分数:[^dsa-kernel]

1
2
3
4
5
6
7
8
9
10
11
12
13
q_fp8, q_scale = act_quant(q, block_size, self.scale_fmt)
k_fp8, k_scale = act_quant(k, block_size, self.scale_fmt)
self.k_cache[:bsz, start_pos:end_pos] = k_fp8
self.k_scale_cache[:bsz, start_pos:end_pos] = k_scale
weights = self.weights_proj(x.float()) * self.n_heads ** -0.5
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
index_score = fp8_index(
q_fp8,
weights,
self.k_cache[:bsz, :end_pos],
self.k_scale_cache[:bsz, :end_pos],
)
topk_indices = index_score.topk(min(self.index_topk, end_pos), dim=-1)[1]

算法目标的 decode 语义可以写成:

1
2
3
4
5
6
I_j = Σ_i w_i ReLU(q_i k_j^T)                         [B,T]
J = TopK(I,k) [B,k]
a_j = (qC' cKV_j^T + qR kR_j^T) / sqrt(d), j∈J [B,H,k]
p = Softmax(a, axis=selected_history) [B,H,k]
u = Σ_(j∈J) p_j cKV_j [B,H,Rkv]
y = WO(Concat_h(u WUV_h)) [B,D]

其中 qC'=qC WUK 是 memory-optimal decode 中的代数吸收对象,不是另一份 cache;WUK/WUV 都来自联合 KV up-projection 的 K/V 分区。固定 Python reference 为验证语义,实际先形成 dense 主 score 再用 topk_indices scatter mask;论文目标的 production sparse kernel 则应先 gather 选中 latent,再做主 MLA。[^dsa-model]

![DSA 双缓存与 top-k 读取逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/d0c32b41ddf32ea71df8c1eebbdb52a6.png){ width=98% }
物理对象图:Indexer cache 和 MLA latent cache 都随历史长度增长;稀疏发生在主 MLA 的高维读取,不代表 Indexer 不扫描历史。

训练时也会形成整段 Indexer K 与 MLA latent,但它们是当前 batch 的 activation,而不是跨请求持久 cache。若沿序列切分,top-k 不能只在本 rank 决定:每个 rank 的 local top-k 只是候选,必须再合并成 global top-k,主 MLA 才能读取正确位置。

GDN:固定状态替代逐 token cache

GDN 把历史写入每个 head 的矩阵状态 S_t∈R^(K×V),不保存每个 token 的 K/V。它先衰减旧状态,再计算旧状态对当前 key 的 value 预测,最后只写入预测误差。[^gdn-paper]

对象 公式 / 源码名 Shape 生产者 → 消费者 生命周期
Q/K/V q_t,k_t,v_t [B,H,K][B,H,K][B,H,V] 投影、ShortConv、SiLU、Q/K L2Norm → recurrence 临时 activation
衰减与写入率 lambda_t,beta_t [B,H] a_proj/b_proj 与参数化 → state update 临时 activation
衰减后状态 Sbar_t [B,H,K,V] lambda_t S_(t-1) → 预测、写入 语义中间量;kernel 可融合
预测与误差 vhat_t,e_t [B,H,V] Sbar_t^T k_tbeta(v-vhat) → delta 写入 临时 activation
递归状态 S_t / h [B,H,K,V];cache 可为 [B,H,V,K] 上一步 → 当前读写 → 下一步 固定大小跨 token state
短卷积状态 conv_state_q/k/v 三个有界窗口 ShortConv 上一步 → 下一步 固定大小 cache
读出 r_t,y_t [B,H,V] → [B,D] q_t^T S_t、Gated RMSNorm、WO 临时 activation

完整单 token 递推为:

1
2
3
4
5
6
Sbar_t = lambda_t S_(t-1)                     [B,H,K,V]
vhat_t = Sbar_t^T k_t [B,H,V]
e_t = beta_t (v_t - vhat_t) [B,H,V]
S_t = Sbar_t + k_t e_t^T [B,H,K,V]
r_t = q_t^T S_t [B,H,V]
y_t = WO(GatedRMSNorm(r_t,z_t)) [B,D]

FLA naive reference 中 hSb_v 先是 v_t,随后被复用为误差 e_t:[^gdn-naive]

1
2
3
4
5
6
7
h = h.clone() * g[:, :, i].exp()[..., None, None]
b_v = v[:, :, i].clone()
b_k = k[:, :, i]
b_v = b_v - (h.clone() * b_k[..., None]).sum(-2)
b_v = b_v * beta[:, :, i][..., None]
h = h.clone() + b_k.unsqueeze(-1) * b_v.unsqueeze(-2)
o[:, :, i] = torch.einsum("bhd,bhdm->bhm", q[:, :, i], h)
![GDN 固定状态在训练与推理中的逻辑](https://pic.shaojiemike.top/shaojiemike/2026/07/82b7f691e5f7d37ff047994d2da1f26d.png){ width=98% }
物理语义图:decode 保存双圈状态 `S_t` 与短卷积窗口;训练的 chunk kernel 会产生边界状态,但没有逐 token 持久 KV cache。

训练为了并行,不会真的用 Python 按 token 慢循环。FLA 使用 chunk recurrence 或 fused kernel,并在需要时输出 final state;短序列 decode 则走 recurrent kernel。两条路径实现同一递推,但物化的中间量和 workspace 不同。[^gdn-layer]

固定大小不等于无损

MHA/DSA 仍为每个历史 token 保存独立表示;GDN 让多个 token 竞争同一状态矩阵容量。它用固定内存换取了历史压缩归纳偏置,不应只凭 cache 大小判断模型质量。

CP:训练时交换临时 K/V 分块

Megatron Context Parallelism 沿 sequence 维切分网络输入和所有 activation。Linear、RMSNorm 等不跨 token,可直接在本地 S/P 个位置上执行;Attention 的本地 Query 仍必须看到全局 K/V。官方语义是前向收集 K/V,反向对 K/V activation gradient 做 reduce-scatter,高性能路径以 P2P Ring 实现。[^mcore-cp]

设 Rank r 初始拥有:

1
Q_r,K_r,V_r : [B,H,S/P,d]

Ring 路径固定 Q_r,让 (K_j,V_j) 依次经过所有 rank。每轮用 online softmax 合并局部分数,避免物化完整 [S/P,S] 分数矩阵:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
def cp_ring_forward(q_local, k_local, v_local, cp_group):
running_max = full([B, H, S_local, 1], -inf)
running_sum = zeros([B, H, S_local, 1])
running_out = zeros([B, H, S_local, head_dim])
k_block, v_block = k_local, v_local
owner = cp_group.rank
for round_id in range(cp_group.size):
scores = matmul(q_local, transpose(k_block, -1, -2)) * scale
scores = apply_global_causal_mask(scores, q_owner=cp_group.rank, kv_owner=owner)
block_max = scores.max(dim=-1, keepdim=True)
new_max = maximum(running_max, block_max)
old_factor = exp(running_max - new_max)
block_exp = exp(scores - new_max)
running_out = running_out * old_factor + matmul(block_exp, v_block)
running_sum = running_sum * old_factor + block_exp.sum(dim=-1, keepdim=True)
running_max = new_max
k_block, v_block, owner = ring_send_recv(k_block, v_block, owner, cp_group)
return running_out / running_sum

这里 K_j,V_j 是当前训练 forward 的临时 activation block。反向还需把各轮产生的 dK_j,dV_j 送回 owner 并累加;它们不会像推理 cache 那样在下一次 forward 继续存在。

![CP 的临时 KV 分块与 Ring 计算](https://pic.shaojiemike.top/shaojiemike/2026/07/84e26ec1e3beaa6d9574dc954be2feca.png){ width=98% }
训练物理语义图:本地 Q 固定,K/V block 轮转,online-softmax 状态留在 Query owner;反向的 KV 梯度回到原 owner。

对于 MQA/GQA,Hkv 更小,Ring 发送的 K/V 体积也更小,NVIDIA 文档明确把它列为降低 CP 通信量的手段。[^mcore-cp] 但这不改变“本地 Q 要覆盖全局 token”的语义。

USP:先换轴,再沿 Ring 扩展上下文

USP 是 Ulysses 与 Ring 的二维组合。令总 sequence-parallel degree 为:

1
P = Pu × Pr

初始每个 rank 拥有 [B,S/(PuPr),H,D]。Ulysses group 中的 All-to-All 把一部分 sequence 分片换成 head 分片:

1
2
3
[B,S/(PuPr),H,D]
-- A2A(seq → head) -->
[B,S/Pr,H/Pu,D]

随后 Ring group 在 Pr 个 rank 之间轮转 K/V,最后 inverse A2A 把输出恢复到初始 layout。YunChang 的公开实现正是按照 process group 初始化、local shard 提取、LongContextAttention 的顺序组织,并要求 world_size=Pu×Pr。[^usp-code]

1
2
3
4
5
6
7
8
9
10
11
def usp_attention(q_local, k_local, v_local, ulysses_group, ring_group):
q_head = all_to_all(q_local, scatter_axis="head", gather_axis="sequence",
group=ulysses_group)
k_head = all_to_all(k_local, scatter_axis="head", gather_axis="sequence",
group=ulysses_group)
v_head = all_to_all(v_local, scatter_axis="head", gather_axis="sequence",
group=ulysses_group)
out_head = ring_online_attention(q_head, k_head, v_head, group=ring_group)
out_local = all_to_all(out_head, scatter_axis="sequence", gather_axis="head",
group=ulysses_group)
return out_local
![USP 的二维换轴与 Ring KV 轮转](https://pic.shaojiemike.top/shaojiemike/2026/07/95f09e03ed7f2d5b8ef9e3f69fa1eccb.png){ width=98% }
训练物理语义图:`Pu` 维用 A2A 将 sequence 换为 head 分片,`Pr` 维再让 K/V block 轮转;最后逆 A2A 恢复每 rank 的原始 sequence shard。

纯 Ulysses 的并行度受可切分 head 数限制;Ring 不切 head,但增加 P2P 轮次。USP 只让 Pu 消耗 head 维,Pr 可继续扩展,因此适合把节点内高带宽交给 A2A、节点间交给 Ring。论文在特定 2×8 A800、LLaMA3-8B 配置下报告 208K 序列的 47% MFU;这只是该硬件和配置的证据,不是普遍性能保证。[^usp-paper]

![USP 论文中的 Ulysses 与 Ring Attention](https://pic.shaojiemike.top/shaojiemike/2026/07/a3d4240389e511a794248dfce06e4807.png){ width=90% }
USP 论文 Figure 2:左侧是 Ulysses 换轴,右侧是 Ring 的 KV block 轮转;USP 把二者放到正交 process group。
![USP 在 208K 序列长度下的性能实验](https://pic.shaojiemike.top/shaojiemike/2026/07/7d5cceaab46c046afeb1a46c8b67da95.png){ width=90% }
USP 论文 Table 7:2×8 A800、LLaMA3-8B 的 TP × Ulysses × Ring 组合。不能把该结果外推到不同网络、head 布局与模型。

训练切分与推理 cache 分片

MHA decode:需要 distributed softmax

训练 CP/USP 的输入有长 sequence 维,可以把 Q/K/V activation 一起切开。decode 通常只有 q_t 一个 Query;若把持久 K/V cache 按历史 token 分给 P 个 rank,每个 rank 只能算局部 score,必须按全局 Softmax 合并。

Rank r 对本地 cache 得到局部最大值 m_r、局部分母 l_r 与未归一化输出 u_r

1
2
3
m_r = max_j a_(r,j)
l_r = Σ_j exp(a_(r,j)-m_r)
u_r = Σ_j exp(a_(r,j)-m_r) v_(r,j)

全局结果是:

1
2
3
m = max_r m_r
l = Σ_r exp(m_r-m) l_r
o = (Σ_r exp(m_r-m) u_r) / l

所以“每卡只放 T/P cache”并不免费:还要复制或切分新 Query、归并 max/sum/output,并处理 PagedAttention 的 block table、请求迁移和负载均衡。只有 cache 容量或带宽收益超过通信成本时才值得。

DSA:先全局选,再稀疏读

DSA 的 cache 若按历史 token 分片,每个 rank 可先为本地 Indexer K 算 local_topk,但随后必须:

  1. 合并所有 rank 的候选分数与全局位置,得到真正的 global_topk
  2. 按 owner gather 选中的 cKV/kR,或把当前 Query 发到 owner 上完成局部主 MLA。
  3. 对选中位置的主 attention 做全局 Softmax 合并。

若 USP 的 Ulysses 维还切分了 Indexer heads,I_j=Σ_i w_i ReLU(q_i k_j^T) 的跨 head 求和也要先 reduce,才能做 top-k。标准 MHA 的 A2A/Ring wrapper 因此不能不加修改地保证 DSA 语义。

GDN:传状态,不传 KV

GDN 没有历史 K/V block 可供 Ring 轮转。若训练序列按连续 chunk 分在多个 rank,后一个 chunk 的第一个 token 需要前一个 chunk 的结束状态:

1
S_end^(r-1) → S_init^r

最直观实现是按 causal 顺序传 boundary state;更并行的实现可以让每个 chunk 形成“输入状态 → 输出状态”的 transition summary,再做 prefix scan。反向传播依赖方向相反。若用 Ulysses 只按 head 切分,GDN 的 head-wise 状态相互独立,通常更自然;若再增加 Ring 序列维,就要把 KV Ring 改成状态传递或 scan。

这三段是从计算依赖推出的实现条件,不表示 Megatron CP 或 YunChang 已经为任意 DSA/GDN 模型提供开箱即用的训练与 serving 支持

联合比较

场景 MHA DSA GDN
训练历史对象 全序列 Q/K/V activation Indexer K + MLA latent activation chunk activation + boundary state
推理持久对象 K_cache,V_cache Indexer K/scale + cKV/kR S_t + 短卷积 state
随上下文 T 增长 是,O(T Hkv d) 是,两套压缩历史 否,相对 T 为固定大小
Decode 主读取 全历史 K/V Indexer 全历史;主 MLA top-k 当前固定状态
CP 序列切分 KV block 全局交互 global top-k + selected latent boundary state / prefix scan
USP 的 Ulysses 维 切 heads 后本地 attention Indexer 跨 head score 可能需 reduce head-wise state 较自然
USP 的 Ring 维 KV block 轮转 Indexer/latent 分片与全局 top-k 不轮转 KV,改传 state

显存峰值应分项记录:

1
2
3
4
5
6
M_peak =
M_model_state
+ M_saved_activation
+ M_persistent_cache_or_state
+ max_t(M_operator_workspace(t) + M_comm_buffer(t))
+ M_allocator_margin

不能只把某个对象除以 P 就声称总显存严格缩小 P 倍。All-to-All 接收 buffer、Ring 双缓冲、FlashAttention workspace、GDN chunk 中间量和 allocator 碎片都可能形成新的峰值。

常见误区

  1. “训练也有 K/V,所以就是 KV cache”:名字相同,生命周期不同。训练 K/V 属于本轮计算图;推理 cache 跨 decode step 存活。
  2. “用了 KV cache,decode 就是 O(1)”:省掉的是历史 K/V 重算;dense MHA 的新 Query 仍读取 T 个历史位置,单步 attention 仍随 T 线性。
  3. “DSA 只保存 top-k”:错误。为了下一个 Query 能重新选择,Indexer 与 MLA 历史表示仍要保存;top-k 是每个 Query 的读取集合。
  4. “GDN 是另一种压缩 KV cache”:更准确地说它是递归状态。旧 token 不再拥有可逐项寻址的独立 K/V。
  5. “CP 就是每卡只看局部上下文”:每卡只保存局部 activation,但本地 Query 的数学结果仍须包含全局允许的 K/V。
  6. “训练 USP 配置可直接照搬到 decode”:decode 的 Query sequence 长度通常为 1,持久 cache 的放置、请求调度和归并通信需要独立设计。
  7. “USP 总并行度只要不超过 Q heads 就行”:真正消耗 head 维的是 Pu,而 GQA/MQA、Indexer 和 TP 共存时还要核对 KV/head 映射与跨 head reduce。

总结

理解 cache 的最好方法不是背缩写,而是对每个对象问四遍:

  1. 谁产生它?
  2. 谁在什么时候再次读取它?
  3. 它活到当前 forward、当前 chunk,还是下一次 decode step?
  4. 沿 sequence 或 head 切分后,哪个归约才能保持与单卡相同的数学结果?

沿这四个问题看,MHA 是逐 token 历史表,DSA 是“全历史轻索引 + top-k 主读取”,GDN 是固定状态机;CP/USP 则是训练时重排全局交互的办法。它们可以组合,但绝不能仅凭都出现了 K、V、state 或 sequence shard,就认为保存与通信语义相同。

参考资料

[^hf-cache]: Hugging Face, Transformers Caching, accessed 2026-07-23.
[^mcore-cp]: NVIDIA, Megatron Core Context Parallel Package, accessed 2026-07-23.
[^usp-paper]: Fang et al., USP: A Unified Sequence Parallelism Approach for Long Context Generative AI, first submitted 2024-05-13.
[^usp-code]: YunChang, hybrid Attention implementation, commit 56118e0d.
[^dsa-release]: DeepSeek-AI, DeepSeek-V3.2-Exp, released 2025-09-29.
[^dsa-model]: DeepSeek-AI, inference/model.py, commit 87e509a2.
[^dsa-kernel]: DeepSeek-AI, inference/kernel.py, commit 87e509a2.
[^gdn-paper]: Yang et al., Gated Delta Networks: Improving Mamba2 with Delta Rule, first submitted 2024-12-09.
[^gdn-layer]: FLA, fla/layers/gated_deltanet.py, commit c70f11c5.
[^gdn-naive]: FLA, fla/ops/gated_delta_rule/naive.py, commit c70f11c5.

Author

Shaojie Tan

Posted on

2026-07-23

Updated on

2026-07-23

Licensed under